Skip to main content

04 - 从单卡到多卡

这一篇把 03 篇的 train.py 改成能在多张卡上跑,产出 train_ddp.py。同时要回答一个问题:多卡到底快了多少,代价是什么。

这也是整个专题里唯一必须真开多卡才能验的一篇。原理和账我会算清楚,但最终那张实测表得在租来的机器上填。

前置:03 篇的单卡训练已经跑通。

零、开始之前:多卡在并行什么

0.1 三种切法

一个模型要放到多张卡上,无非是切三样东西:

叫法切什么直觉
数据并行切数据每张卡放一份完整模型,各自算不同的数据,再把梯度对齐
张量并行切单个矩阵一个大矩阵乘法拆到多张卡上算,算完拼起来
流水线并行切层前 9 层在卡 0,后 9 层在卡 1,像流水线一样传递

真正的大模型训练是三种混着用,叫 3D 并行。

一个模型放到多张卡上,无非切三样东西① 数据并行 —— 切数据整个模型数据 A整个模型数据 B整个模型数据 C整个模型数据 D每张卡一份完整模型,各算不同的数据,再把梯度对齐 —— 这一篇只讲它我们用这�个② 张量并行 —— 切单个矩阵卡0卡1卡2卡3同一次大矩阵乘拆到多张卡上算,算完拼起来为「连单层都装不下」而生用不上③ 流水线并行 —— 切层层 1-4层 5-910-1415-18前 9 层在卡 0、后 9 层在卡 1,像流水线传递同样为「装不下」而生用不上0.5B 单卡放得下(02 篇算过静态 8.04 GB),所以 ② 和 ③ 完全用不上 —— 它们存在的前提是「一张卡装不下一个模型」,我们没这个问题。真正的大模型训练是三种混着用,叫 3D 并行。这一篇只讲 ①,但会讲透:怎么工作、通信量多大、什么时候不够用、换成 FSDP 又要付什么代价。
先分清这三者,后面的取舍才有依据。一个常见的误解是「多卡 = 模型被拆开了」,但数据并行里每张卡装的都是完整模型,被拆的只有数据。也正因为如此,它是三者里唯一不改变模型计算方式的一种。

0.2 我们只用数据并行

0.5B 的模型单卡放得下(02 篇算过静态显存 8 GB),所以张量并行和流水线并行完全用不上。它们是为了「一张卡装不下一个模型」而存在的,我们没这个问题。

这一篇只讲数据并行,但会讲透:DDP 怎么工作、通信量多大、什么时候它不够用、不够用了换成 FSDP 又要付什么代价。

0.3 数据并行的基本流程

每张卡各拿一批不同的数据,各自跑前向反向,得到各自的梯度。这些梯度不一样,因为看的数据不一样。

关键在于:参数更新前必须让所有卡的梯度变成同一个值,否则各卡的参数就会越走越远,最后变成四个不同的模型。

对齐的方式是求平均。这个「把所有卡的数据加起来再分回去」的操作,就是下一节的 all-reduce。

梯度不一样是正常的,要对齐的是「更新」这一步卡0 数据 A → 梯度 g0卡1 数据 B → 梯度 g1卡2 数据 C → 梯度 g2卡3 数据 D → 梯度 g3all-reduce 求平均四份梯度加起来除以 4实现 = reduce-scatter + all-gather卡0 平均梯度 → step()卡1 平均梯度 → step()卡2 平均梯度 → step()卡3 平均梯度 → step()四张卡拿到的是同一个梯度,参数从此一步不差关键约束:参数更新之前,必须让所有卡的梯度变成同一个值。少了这一步不报错,只是四张卡各走各的,训完得到四个略微不同的模型。注意梯度不同并不是错误 —— 每张卡看的数据本来就不一样。数据并行的全部技术含量,就在中间那一格怎么做得又快又不成为瓶颈。
这张图解释了为什么数据并行的加速比永远拿不满:中间那一格是四张卡的同步点,任何一张卡慢了,其余三张都得等。后面 2.3 节算的通信量,量的就是这一格的成本。

一、集合通信

1.1 几个原语

多卡之间怎么交换数据,是一组标准操作,叫集合通信(collective communication)。NVIDIA 的实现叫 NCCL,PyTorch 底层用的就是它。

操作干什么
broadcast把某张卡的数据复制给所有卡
reduce把所有卡的数据求和(或求最大等),结果放在一张卡上
all-reduce求和,结果每张卡都有一份
reduce-scatter求和,但每张卡只拿走结果的一部分
all-gather每张卡的一小块,拼成完整的一份,每张卡都有

要记的关键关系是:all-reduce = reduce-scatter + all-gather。这个拆解决定了它的通信量,也是 FSDP 的实现基础。

五个原语,看「四张卡上的内容前后变成什么」就够了调用前(卡0 · 卡1 · 卡2 · 卡3)调用后broadcastAAAAA把一张卡的数据复制给所有卡reduceabcdΣ求和,结果只落在一张卡上all-reduceabcdΣΣΣΣ求和,每张卡都有一份 —— DDP 用它reduce-scatter全量全量全量全量Σ块0Σ块1Σ块2Σ块3求和,但每张卡只拿走结果的一块all-gather块0块1块2块3全部全部全部全部每张卡的一块拼成完整的一份要记的关键关系:all-reduce = reduce-scatter + all-gather。这个拆解既决定了 ring all-reduce 的通信量系数 2×(N−1)/N,也是 FSDP 的实现基础。
把最后两行和第三行对着看:all-reduce 之所以能做到「每张卡都拿到完整结果」而不让任何一张卡成为流量瓶颈,靠的就是先各拿一块求和、再互相补齐。FSDP 用的是同一套零件,只是把它们拆开分别用在参数和梯度上。

1.2 ring all-reduce 的通信量

朴素的做法是所有卡都把数据发给卡 0,卡 0 加完再发回去。这样卡 0 的网卡要承担 N 倍流量,成为瓶颈。

实际用的是 ring all-reduce:把 N 张卡连成一个环,数据切成 N 块。

数据切成 N 块,分两阶段各转 N-1 步:

4 张卡连成环,梯度切成 4 块卡 0卡 1卡 2卡 3每步只跟右邻居收发一块起始:每张卡各有全部 4 块,但都只是自己那份块0块1块2块3卡0卡1卡2卡3阶段一结束:每卡各持有一块的完整和块0块1块2块3卡0卡1卡2卡3已是 4 卡求和结果还只是局部值阶段一 · reduce-scatter,转 3 步每步把手里的一块发给右邻居,同时收左邻居的一块并加上去。3 步之后,卡 i 手里的第 i 块已经累加了全部 4 张卡的值。每卡收发 (N−1)/N × 数据量阶段二 · all-gather,再转 3 步把各自那块完整结果沿环传一圈,不再做加法,只是复制。3 步之后每张卡都集齐了 4 块完整结果。每卡收发 (N−1)/N × 数据量合计每卡收发 2 × (N−1)/N × 数据量。N 越大这个系数越接近 2,与卡数无关 —— 这是它能扩展到几千卡的原因。对照:朴素做法是所有卡都发给卡 0、卡 0 加完再发回,卡 0 的网卡要承担 N 倍流量,直接成为瓶颈。本专题 0.5B 模型 4 卡:梯度 1.004 GB × 1.5 = 每步每卡收发 1.51 GB。
橙色块表示「只有本卡的局部值」,绿色块表示「已经是四卡求和的完整结果」。第一阶段让每张卡各自凑齐一块,第二阶段再把这些块散播回去。两阶段各转 N−1 步,任何时刻每条链路上只跑一块数据,带宽被均匀利用。

第一阶段 reduce-scatter,转 N-1 步,每步每张卡收发一块,结束后每张卡持有完整结果的其中一块。第二阶段 all-gather,再转 N-1 步,把每块广播到所有卡。

所以每张卡收发的数据量是:

通信量=2×N1N×数据总量\text{通信量} = 2 \times \frac{N-1}{N} \times \text{数据总量}

N 很大时这个系数趋近 2,跟卡数无关。这是 ring all-reduce 的核心优点,也是它能扩展到几千张卡的原因。

二、DDP

2.1 原理

PyTorch 的 DistributedDataParallel(DDP)做的事:

  1. 启动时把 rank 0 的参数 broadcast 给所有卡,保证起点一致
  2. 每张卡各自前向反向
  3. 反向过程中,梯度一算出来就 all-reduce,跟反向计算重叠
  4. 每张卡拿到平均梯度,各自 optimizer.step()

第 3 点是 DDP 快的关键,也是它和朴素实现的主要差距:

朴素实现:算完全部反向,再统一通信计算层18层17层1通信all-reduce 全部梯度一步总时长DDP:分桶,某桶梯度一齐就开传计算层18层17层1通信桶1桶2桶3桶4一�步总时长(短得多)反向是从最后一层往前算,所以最后一层的梯度最先就绪,可以立刻开始传。重叠只能藏住通信,藏不掉通信量。真正减少通信量的是下一节的 no_sync —— 那一项能砍掉 16 倍。
反向传播从最后一层往前算,所以最后一层的梯度最先就绪。DDP 把参数分成若干桶,某个桶凑齐就立刻发出,与仍在进行的反向计算并行。通信因此被藏在计算的阴影里,但注意它减少的是等待时间,不是数据量。

注意每张卡都是完整地跑一遍 optimizer.step(),输入(平均梯度)相同、初值相同,所以结果必然相同。参数不需要额外同步。

2.2 代码改动

03 篇的 train.py 改成多卡,实际只动四个地方。

一是初始化进程组、绑定设备。

def setup_dist():
dist.init_process_group(backend="nccl")
rank = int(os.environ["RANK"])
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(local_rank) # 必须在建模型之前
return rank, local_rank, world_size, torch.device(f"cuda:{local_rank}")

RANKLOCAL_RANKWORLD_SIZE 这三个环境变量由 torchrun 注入,不用自己管。rank 是全局编号,local_rank 是这台机器内的编号,单机情况下两者相同。

torch.cuda.set_device(local_rank) 必须在建模型之前调用,否则所有进程会往卡 0 上堆,直接 OOM。

二是包一层 DDP。

model = DDP(model, device_ids=[local_rank])

三是随机种子要错开。

torch.manual_seed(args.seed + rank)
np.random.seed(args.seed + rank)

我们的 get_batch 是随机取起点的,如果所有卡用同一个种子,四张卡会取到完全一样的数据,等于白算三份。加上 rank 就错开了。

四是日志和存盘只在 rank 0 做。

is_main = rank == 0
if is_main and step % args.log_every == 0:
print(...)

不然四个进程会往屏幕上刷四份一样的日志,checkpoint 也会被互相覆盖。

2.3 通信量算一算

00-pretrain-0.5b/parallel_math.py 的实际输出:

二、每卡每步通信量
方案 通信量 NVLink PCIe
DDP + no_sync(每步只同步一次) 1.51 GB 2.5ms 60.3ms
DDP 无 no_sync(每次累积都同步) 24.11 GB 40.2ms 964.2ms
FSDP full shard 36.16 GB 60.3ms 1446.3ms

三、通信占一步的比例(一步约 6.42 秒)
方案 NVLink PCIe
DDP + no_sync(每步只同步一次) 0.04% 0.94%
DDP 无 no_sync(每次累积都同步) 0.63% 15.02%
FSDP full shard 0.94% 22.53%

先看第一行。DDP 每步要 all-reduce 一遍全部梯度,bf16 梯度是 1.004 GB,乘上 ring 的系数 2×34=1.52 \times \frac{3}{4} = 1.5,得到 1.51 GB。

在 NVLink 上这只要 2.5 ms,占一步(6.42 秒)的 0.04%,基本免费。在 PCIe 上是 60 ms,占 0.94%,也还好。

结论是:0.5B 模型配 16 的梯度累积,DDP 的通信开销几乎可以忽略。 这个结论有点反直觉,很多人以为多卡一定被通信拖累。原因下一节说。

每卡每步通信量(一步约 6.42 秒)DDP + no_sync每步只同步一次1.51 GBNVLink 2.5 ms · 0.04%PCIe  60.3 ms · 0.94%DDP 无 no_sync每次累积都同步24.11 GBNVLink 40.2 ms · 0.63%PCIe  964.2 ms · 15.02%FSDP full shard前向反向各 all-gather 一遍36.16 GBNVLink 60.3 ms · 0.94%PCIe  1446.3 ms · 22.53%反直觉的结论:0.5B 配 grad_accum=16,DDP 在 NVLink 上的通信只占 0.04%,基本免费 —— 很多人以为多卡一定被通信拖累,在这个规模上并不成立。但 PCIe 那一行完全是另一回事:FSDP 有近四分之一的时间在等通信。开卡后第一件事是 nvidia-smi topo -m,看清楚自己该读哪一列。
三根条形的长度比就是 1 : 16 : 24。第一根之所以这么短,全靠 no_sync 把 16 次同步压成 1 次;第三根之所以这么长,是因为 FSDP 每个 micro batch 都要重新 all-gather 参数,压根没有可以摊薄的余地。互联方式在这里的放大倍数是 24 倍,选机器时值得先看这一项。

2.4 no_sync:最重要的一个优化

注意上表第二行,不用 no_sync 的话通信量是 24.11 GB,整整 16 倍

原因在于 DDP 的默认行为是「每次 backward() 都触发 all-reduce」。而我们有 16 次梯度累积,就是 16 次 backward(),于是同步了 16 遍。但前 15 次的梯度还没累加完,同步它们毫无意义。

no_sync() 这个上下文管理器就是用来关掉中间那些同步的:

for micro in range(args.grad_accum):
x, y = get_batch(train_data, args.micro_batch, args.seq_len, device)

last = (micro == args.grad_accum - 1)
sync_ctx = nullcontext() if last else model.no_sync()
with sync_ctx:
with ctx:
_, loss = model(x, y)
loss = loss / args.grad_accum
loss.backward()

只有最后一次累积走正常路径触发 all-reduce,前 15 次都在 no_sync() 里,梯度只在本地累加。

这就解释了 2.3 节那个反直觉的结论:梯度累积把通信摊薄了。累积次数越多,同一次通信服务的计算越多,通信占比越低。反过来说,如果 grad_accum=1,通信占比会直接乘 16。

这行代码少写不会报错,loss 曲线也完全正常,只是白白多花十几倍的通信。在 PCIe 机器上这个差别是 0.94% 对 15.02%,很可观。

一个训练步内有 16 次 backward,同步几次由你决定计算 backward通信 all-reduce❌ 默认行为16 次 backward → 16 次 all-reduce → 24.11 GB / 步。前 15 次的梯度根本还没累加完,同步它们毫无意义。✅ 加no_sync前 15 次都在 no_sync() 里,梯度只在本地累加;只有最后一次走正常路径触发同步 → 1.51 GB / 步,整整少 16 倍。这也解释了那个反直觉的结论:梯度累积把通信摊薄了。累积次数越多,同一次通信服务的计算越多,通信占比越低;grad_accum=1 时通信占比直接乘 16。这行代码少写不会��报错,loss 曲线也完全正常,只是白花十几倍的通信。NVLink 上看不太出来,PCIe 上是 15.02% 对 0.94%。
红色细条的总面积就是账单。值得注意的是两行的计算部分完全一样 —— 也就是说这个优化不改变任何数值结果,纯粹是把无意义的同步删掉。这类「不影响正确性、只影响账单」的开关,是多卡训练里最值得先排查的一类。

2.5 用 torchrun 启动

torchrun --standalone --nproc_per_node=4 train_ddp.py --mode ddp

--nproc_per_node=4 表示起 4 个进程,一个进程一张卡。--standalone 是单机模式,多机要改成指定 master 地址和端口。

不要用 python train_ddp.py,那样只会起一个进程,环境变量也不会注入。

三、显存不够的时候:ZeRO 与 FSDP

3.1 DDP 浪费在哪

DDP 每张卡都存一份完整的参数、梯度、优化器状态。四张卡就是四份一模一样的优化器状态。

02 篇算过每参数 16 字节,其中 12 字节是优化器状态(fp32 master weights 加 Adam 的两个动量)。这 12 字节在四张卡上完全重复,是纯浪费。

ZeRO(Zero Redundancy Optimizer)要解决的就是这个。

3.2 ZeRO 的三个阶段

按切分程度分三级,一级比一级省:

阶段切什么每参数字节(4 卡)
ZeRO-0(就是 DDP)什么都不切16
ZeRO-1切优化器状态7
ZeRO-2再切梯度5.5
ZeRO-3再切参数本身4

ZeRO-3 就是 PyTorch FSDP 的 FULL_SHARD 模式。

四个阶段各自切了什么,画出来:

参数 2 字节梯度 2 字节优化器状态 12 字节(fp32 主权重 + Adam 双动量)色块宽度按字节数等比绘制,每格 = 一张卡持有的份额ZeRO-0(即 DDP):三样都完整复制 4 份卡0卡1卡2卡316 字节 / 参数 · 8.04 GB / 卡ZeRO-1:只切优化器状态卡0卡1卡2卡37 字节 / 参数 · 3.52 GB / 卡ZeRO-2:再切梯度卡0卡1卡2卡35.5 字节 / 参数 · 2.76 GB / 卡ZeRO-3 / FSDP:参数也切卡0卡1卡2卡34 字节 / 参数 · 2.01 GB / 卡切得越多,�省的是显存16 → 7 → 5.5 → 4 字节/参数,8.04 GB 压到 2.01 GB。ZeRO-1 只切优化器状态就砍掉一半多。花的是通信ZeRO-3 平时只存分片,用到某层才 all-gather 出完整参数,算完立刻扔 —— 通信量涨到 DDP 的 24 倍。
色块宽度按每参数字节数等比绘制,一眼能看出优化器状态(琥珀色,12 字节)才是大头,这也是 ZeRO-1 只切它就能从 16 降到 7 的原因。左下和右下两块要连起来读:显存曲线和通信曲线是反向的,0.5B 这个规模用 ZeRO-3 是纯亏 —— 显存本来就富余,白付 24 倍通信。

ZeRO-3 的思路:平时每张卡只存 1/N 的参数,要用到某一层的时候临时 all-gather 出完整参数,算完立刻扔掉。这样显存里任何时刻只有一层的完整参数,其余都是分片。

3.3 显存账

00-pretrain-0.5b/parallel_math.py 的输出:

一、每卡静态显存(不含激活值)
方案 字节/参数 显存
DDP / ZeRO-0(什么都不切) 16.00 8.04 GB
ZeRO-1(切优化器状态) 7.00 3.52 GB
ZeRO-2(再切梯度) 5.50 2.76 GB
ZeRO-3 / FSDP full shard(全切) 4.00 2.01 GB

从 8.04 GB 降到 2.01 GB,省了四倍,正好等于卡数。

但对我们这个模型,这个收益毫无意义,因为 8.04 GB 本来就装得下,省到 2 GB 只是把用不上的显存空出来。

3.4 FSDP 的代价

省显存不是白省的。看 2.3 节那张通信表的第三行:FSDP 是 36.16 GB,比 DDP 加 no_sync24 倍

为什么这么多。FSDP 每次前向要 all-gather 一遍参数,反向还要再 all-gather 一遍(因为前向算完就扔了),梯度还要 reduce-scatter。三趟,而且每个 micro batch 都要走一遍,没法像 DDP 那样靠 no_sync 摊薄。

把 DDP 和 FSDP 的一步画在一起,这三趟就很显眼了:

标准数据并行与完全分片数据并行的一步对照:上半部分 DDP 的每张卡是 model shard、forward、backward、all-reduce、update weights,通信只有中间一次 all-reduce 同步梯度;下半部分 FSDP 在 forward 前有一次 all-gather 收齐参数、backward 前再 all-gather 一次、算完梯度还要 reduce-scatter,而且前两者都标着 N LAYERS 表示逐层各来一遍
出处:NVIDIA/Megatron-LM docs/images/megatron_fsdp/DDP_vs_FSDP.png,BSD-3。看下半部分那两个标着 N LAYERS 的灰框 —— all-gather 不是整个模型来一次,而是每一层各来一次。这也是为什么 FSDP 的通信量不只是「多几趟」,而是跟层数挂钩。
16×(2×参数×N1N+梯度×N1N)=36.16 GB16 \times (2 \times \text{参数} \times \tfrac{N-1}{N} + \text{梯度} \times \tfrac{N-1}{N}) = 36.16\ \text{GB}

在 NVLink 上还好(0.94%),在 PCIe 机器上是 22.53%,也就是四分之一的时间在等通信

3.5 什么时候该用哪个

情况用什么
模型 + 优化器状态单卡装得下DDP,别想别的
装不下,但差得不多ZeRO-1 或 ZeRO-2,通信代价小
差很多ZeRO-3 / FSDP full shard
单层都装不下得上张量并行,超出这一篇范围

判据很简单:先按每参数 16 字节算一遍(02 篇 8.5 节),能装下就用 DDP。

判据只有一条:先按每参数 16 字节算一遍,能装下就用 DDP显存缺口:小 → 大装得下DDP别想别的通信 1.51 GB / 步NVLink 上占 0.04%装不下,但差得不多ZeRO-1 / ZeRO-2切优化器状态,再切梯度16 → 7 → 5.5 字节 / 参数通信代价还很小差很多ZeRO-3 / FSDP参数也切,降到 4 字节通信涨到 24 倍PCIe 上占 22.53%连单层都装不下张量并行把单个矩阵乘拆到多卡超出这一篇的范围我们停在第一格:8.04 GB 单卡本来就富余,切到 2.01 GB 只是把用不上的显存空出来,却要付 24 倍通信 —— 纯亏。所以这一篇要人为把模型放大到 4B,才有场景演示 FSDP。顺序永远是先算显存缺口,再谈并行策略,不要反过来。
这张图最该带走的是那条横轴:并行方案不是「越先进越好」的排序,而是按显存缺口大小排的阶梯,每往右一格都用确定的通信代价换确定的显存。缺口是 0 的时候往右走,换来的只有账单。

3.6 FSDP 代码

policy = functools.partial(transformer_auto_wrap_policy,
transformer_layer_cls={Block})
model = FSDP(
model,
auto_wrap_policy=policy,
sharding_strategy=ShardingStrategy.FULL_SHARD,
mixed_precision=MixedPrecision(param_dtype=torch.bfloat16,
reduce_dtype=torch.bfloat16,
buffer_dtype=torch.bfloat16),
device_id=local_rank,
use_orig_params=True,
)

几个要点。

auto_wrap_policy 决定按什么粒度切。按 Block 切是标准做法:切得太细通信次数多,太粗则一次 all-gather 出来的参数太大,省不下显存。

mixed_precision 是 FSDP 自己的混合精度配置,配了它就不要再套 autocast,两套机制会打架。这是 train_ddp.pyctx 只在 DDP 模式下启用的原因。

use_orig_params=True 让参数保持原来的结构,torch.compile 和 03 篇按维度分组的 weight decay 才能正常工作。

梯度裁剪也要换:FSDP 的参数是分片的,普通 clip_grad_norm_ 算出来的范数是错的,得用 model.clip_grad_norm_()

四、动手实测

原理和账都算完了,剩下的必须真跑。

4.1 要测的四组

# 1 单卡基线
python train.py --steps 50

# 2 DDP 四卡
torchrun --standalone --nproc_per_node=4 train_ddp.py --mode ddp --steps 50

# 3 FSDP 四卡
torchrun --standalone --nproc_per_node=4 train_ddp.py --mode fsdp --steps 50

# 4 DDP 四卡但故意去掉 no_sync(改一行代码),验证 2.4 节那 16 倍
torchrun --standalone --nproc_per_node=4 train_ddp.py --mode ddp --steps 50

--steps 50 是只跑 50 步就退出,够测速度了,不用等训练跑完。前几步要跳过(有 CUDA 初始化和编译开销),从第 10 步开始计时。

4.2 待填的实测表

配置tokens/s相对单卡MFU峰值显存
单卡待测1.00×待测待测
DDP 4 卡待测待测待测待测
DDP 4 卡(无 no_sync)待测待测待测待测
FSDP 4 卡待测待测待测待测

按前面的账,预期是这样:DDP 四卡的加速比应该接近 3.7 到 3.9 倍(不会满 4 倍,有通信和同步开销);FSDP 峰值显存明显更低但速度略慢;去掉 no_sync 在 NVLink 机器上差别很小,在 PCIe 机器上会明显掉速。

如果实测和预期差很多,那说明账算错了或者哪里有瓶颈,这比数字本身更有价值。 05 篇会把真实结果和这里的预期做对照。

4.3 故意撑爆单卡

00 篇说过,0.5B 上 FSDP 看不出好处,所以要人为制造一个装不下的场景。parallel_math.py 算了阈值:

四、多大的模型才会撑爆单卡
显卡 DDP 能放下的上限
A100 40G 1.75 B 参数
A100 80G 3.50 B 参数
H100 80G 3.50 B 参数

(按静态 16 字节每参数、留三成显存给激活值估的。)

所以把配置放大到 4B 左右,在 80G 卡上单卡和 DDP 都会 OOM,而 FSDP 能跑起来。比如:

Config(d_model=3072, n_layers=32, n_heads=24, n_kv_heads=8, ffn_dim=8192)

这个实验值得做一次。 亲眼看着 DDP 报 OOM、换成 FSDP 就跑起来,比读十遍 ZeRO 论文管用。跑通即可,不用真训。

多大的模型才会撑爆单卡(静态 16 字节/参数,留三成显存给激活值)我们的 0.5B · 8.04 GB实验放大到约 4BA100 40GDDP 放得下超出,必须切A100 / H100 80GDDP 放得下超出00.5B1.75B3.5B4B0.5B 离两条线都很远,这就是为什么必须人为放大才看得出 FSDP 的价值 —— 不放大,那一节就只能纸上谈兵。放大配置:d_model=3072, n_layers=32, n_heads=24, n_kv_heads=8, ffn_dim=8192。跑通即可,不用真训。亲眼看着 DDP 报 OOM、换 FSDP 就起来,比读十遍 ZeRO 论文管用。
这两条线也是以后估算的现成工具:拿到一个新模型,先用「参数量 × 16 字节」比一下卡的容量,就知道该不该考虑分片。注意上限跟卡数无关 —— DDP 每张卡都存一份完整状态,加卡不会让单卡的上限变高。

五、租卡实操

5.1 开卡之前

01 到 03 篇反复强调过,这里汇总一次:

  • train.bin 已经生成好并传到持久化云盘(01 篇)
  • 单卡训练已经跑通,loss 会降(03 篇)
  • checkpoint 保存和恢复已经验证过(03 篇第十节第 5 条)
  • nvidia-smi topo -m 看一眼卡之间是 NVLink 还是 PCIe,这直接决定 2.3 节那张表看哪一列

5.2 checkpoint 在多卡下的坑

DDP 的 model.state_dict() 会带上 module. 前缀,因为模型被包了一层。存之前要脱掉,否则单卡加载时对不上:

raw = model.module if hasattr(model, "module") else model
torch.save({"model": raw.state_dict(), ...}, path)

FSDP 更麻烦,参数是分片的,直接 state_dict() 拿到的是分片。要用 FullStateDictConfig 把完整参数聚到 rank 0 再存。

还有一点:只让 rank 0 存盘,但所有 rank 都要走到 dist.barrier()。如果只有 rank 0 在存盘、其他 rank 继续往下跑,进程之间会失去同步,后面的 all-reduce 会挂住。

5.3 训练挂了怎么排查

多卡训练卡死是常见现象,通常没有报错,就是不动了。排查顺序:

  1. nvidia-smi 看是不是所有卡的利用率都是 100% 且不变,那多半是某个 collective 在等一个永远不来的对端
  2. 检查是不是有条件分支导致某些 rank 调了 all-reduce 而另一些没调,这是最常见的死因
  3. NCCL_DEBUG=INFO 看 NCCL 的日志
  4. TORCH_NCCL_BLOCKING_WAIT=1,让 NCCL 超时报错而不是无限等

第 2 条特别值得注意。比如「loss 是 NaN 就跳过这一步」这种逻辑,如果只有一张卡的 loss 是 NaN,它跳过了而其他三张卡还在等它 all-reduce,整个训练就死锁了。多卡代码里,所有 rank 必须执行完全相同的 collective 序列。

最常见的死因:rank 之间执行了不同的 collective 序列「loss 是 NaN 就跳过这一步」这段逻辑放到 4 卡上会怎样卡0前向 · loss=NaNcontinue 跳过直接进入下一步卡1前向 · loss 正常backwardall-reduce:等卡 0 …… 永远等不到卡2前向 · loss 正常backwardall-reduce:等卡 0 …… 永远等不到卡3前向 · loss 正常backwardall-reduce:等卡 0 …… 永远等不到没有异常、没有超时、进程都还活着 —— nvidia-smi 看到的是四张卡利用率 100% 且一动不动卡死了按这个顺序查① nvidia-smi 看是不是所有卡都 100% 且不变 —— 多半是某个 collective 在等一个永远不来的对端② 检查有没有条件分支让某些 rank 调了 collective 而另一些没调(最常见)③ NCCL_DEBUG=INFO 看日志  ④ TORCH_NCCL_BLOCKING_WAIT=1,让它超时报错而不是无限等
记住一条硬规则就能避开这一整类问题:多卡代码里,所有 rank 必须执行完全相同的 collective 序列。任何依赖数据内容的分支(跳过 NaN、按 loss 提前退出、只在某些 rank 上做评估)都可能破坏它,而破坏的表现永远是「卡住」而不是「报错」。

六、验收

  • 四张卡都在跑,nvidia-smi 看得到四个进程
  • rank 0 的日志正常,没有四份重复输出
  • DDP 四卡的 loss 曲线和单卡基本重合(同样的总 batch)
  • 加速比测出来了,和 4.2 节的预期对照过
  • no_sync 去掉前后的差异测出来了
  • FSDP 能跑,峰值显存明显低于 DDP
  • 放大到 4B,DDP OOM 而 FSDP 能跑
  • checkpoint 存了能被单卡加载(前缀脱干净了)
  • 主动 kill 一次,四卡能从 checkpoint 恢复

第 3 条是正确性的核心检查。多卡只是加速手段,不该改变训练结果。 如果 DDP 四卡的 loss 曲线和单卡差得明显,说明有 bug,最常见的是种子没错开(四张卡在算同样的数据)或者 loss 没除以累积次数。

七、常见问题

现象大概率是什么原因
启动就 OOM,且只有卡 0 满忘了 torch.cuda.set_device(local_rank)
四份重复日志没判断 rank == 0
加速比只有 2 倍左右数据加载成瓶颈;或没开 no_sync(PCIe 机器)
loss 曲线和单卡差很多种子没按 rank 错开,四卡在算同样的数据
训练卡死不动也不报错某些 rank 少调了一次 collective,看 5.3 节
checkpoint 单卡加载报 key 不匹配module. 前缀没脱
FSDP 下梯度裁剪的值明显不对用了 torch.nn.utils.clip_grad_norm_,要换成 model.clip_grad_norm_
FSDP 报混合精度相关的错autocast 和 FSDP 的 MixedPrecision 同时开了
torchrun 起不来,端口被占上一次训练的僵尸进程还在,pkill -f train_ddp

先在两张卡上验,再上四张

--nproc_per_node=2 能暴露绝大多数多卡 bug(前缀、种子、collective 不对称),而且便宜一半。两卡跑通了再加到四卡,基本不会再有新问题。